Repository navigation
Scale SGDW parameters by (1 - lr * weight_decay) - #554
Open
shaneraphel wants to merge 1 commit into
Open
shaneraphel wants to merge 1 commit into
shaneraphel wants to merge 1 commit into
Conversation
The decay step added -lr * weight_decay to every coordinate, so a parameter of 4 with lr 0.1 and weight decay 0.2 moved to 3.98. The decoupled update in Loshchilov and Hutter (arXiv 1711.05101) scales the parameter by (1 - lr * weight_decay) and then subtracts the step, which lands on 3.92. With a gradient of 1.5 the same step is 3.77. Zero weight decay still subtracts lr times the gradient.
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Fixes #454 and the same mismatch reported in #433. Decoupled weight decay, as written by Loshchilov and Hutter (arXiv 1711.05101), scales the parameter and then subtracts the step:
θ ← (1 − λ μ) θ
θ ← θ − μ m
SGDW.stepdid the gradient step and then calledp.add_(weight_decay, alpha=-lr), which subtracts the constant λμ from every coordinate. On a parameter of 4, with learning rate 0.1, weight decay 0.2 and a zero gradient, that lands on 3.98. The paper scale lands on 3.92. With a gradient of 1.5 the paper order (scale, then subtract) lands on 3.77. The same library's Lion, QHAdam, AdaBelief and Sophia already scale by(1 - lr * weight_decay). SGDW was the one that added a constant, which is also why it trained worse than SGD at the same hyperparameters.The decay now happens before the step, and only when weight decay is non-zero. A zero weight decay still subtracts
lrtimes the gradient, including the first momentum step.Test plan
pytest tests/test_sgdw.py— zero gradient scales[4, -2]to[3.92, -1.96]; gradient 1.5 on 4 lands on 3.77; weight decay 0 with momentum 0.9 matches a plain SGD step.